# Script to import and analyse muon data
import pandas as pd
import numpy as np
import seaborn as sns
import os
import matplotlib.dates as mdates
import matplotlib.pyplot as plt
from scipy.stats import linregress

# ---------------------------------------------------
# LOAD DATA
# ---------------------------------------------------

cols = [
    "time",
    "det1",
    "det2",
    "coincidence",
    "temp",
    "pressure",
    "neutrons"
]

spots = [
    "year",
    "month",
    "day",
    "decdate",
    "estspots",
    "STD",
    "stat1",
    "stat2"
]

df_muons = pd.read_csv(
    "../MuonData/pico_data.txt",
    names=cols,
    parse_dates=["time"]
)

df_spots = pd.read_csv(
    "../MuonData/Sunspot_estimate_2026.txt",
    sep=r'\s+',
    names=spots
)

output_dir = "/home/roymogg/muon_pictures"
os.makedirs(output_dir, exist_ok=True)

#
# 
# ---------------------------------------------------
# CLEAN AND ADJUST MUON DATA
# ---------------------------------------------------

df_spots['Date'] = pd.to_datetime(df_spots[['year', 'month', 'day']])

df_spots["solarflux"] = (df_spots["estspots"] + 73.21) / 1.14

df_muons.loc[df_muons["temp"] > 60, "temp"] = np.nan

df_muons = df_muons.dropna()

# ---------------------------------------------------
# MOVING AVERAGES
# ---------------------------------------------------

df_muons["rate_ma_24h"] = df_muons["coincidence"].rolling(24).mean()
df_muons["pressure_ma_24h"] = df_muons["pressure"].rolling(24).mean()
df_muons["neutrons_ma_24h"] = df_muons["neutrons"].rolling(24).mean()

# ---------------------------------------------------
# BASIC STATISTICS
# ---------------------------------------------------

print("\n=== BASIC STATISTICS ===")
print(df_muons.describe())

# Poisson uncertainty estimate
df_muons["rate_err"] = np.sqrt(df_muons["coincidence"])

# Basic Histogram plot

plt.figure(figsize=(8,6))
plt.hist(df_muons["coincidence"], bins=30)

plt.xlabel("Coincidence Counts")
plt.ylabel("Frequency")
plt.title("Muon Count Distribution")
plt.grid(True)

# ---------------------------------------------------
# PRESSURE and NEUTRON CORRELATION
# ---------------------------------------------------

slope, intercept, r_value, p_value, std_err = linregress(
    df_muons["pressure"],
    df_muons["neutrons"]
)

print("\n=== NEUTRON to PRESSURE CORRELATION ===")
print(f"Slope: {slope:.4f}")
print(f"Correlation coefficient r: {r_value:.3f}")
print(f"p-value: {p_value:.3e}")

plt.figure() # Start another fresh figure
plt.scatter(df_muons["pressure"], df_muons["neutrons"], label="Data", color="blue")
trendline = slope * df_muons["pressure"] + intercept
plt.plot(df_muons["pressure"], trendline, color="red", label=f"Fit (r={r_value:.2f})")
plt.title("Neutron vs Pressure")
plt.xlabel("Pressure")
plt.ylabel("Neutrons")
plt.legend()

file_path = os.path.join(output_dir, "neutron_plot.png")
plt.savefig(file_path, dpi=300, bbox_inches="tight")
plt.show(block=False)
plt.pause(1.0)
#plt.clf()


# ---------------------------------------------------
# PLOTS
# ---------------------------------------------------

# SUBPLOT 1: Rates vs Time
# 1. Define the grid size (e.g., 2 rows, 2 columns) and figure size
fig, axs = plt.subplots(2, 2, figsize=(14, 6))

# 2. Target the specific subplot at row 0, column 0 using axs[0, 0]
axs[0, 0].plot(df_muons["time"], df_muons["coincidence"], alpha=0.5, label="Raw Rate")
axs[0, 0].plot(df_muons["time"], df_muons["rate_ma_24h"], linewidth=2, label="24h Moving Average")
axs[0, 0].xaxis.set_major_locator(mdates.WeekdayLocator(interval=1))
#axs[0, 0].set_xlabel("Time",fontsize=12)
axs[0, 0].set_ylabel("Muon Rate", fontsize=12)
axs[0, 0].set_title("Muon Rate vs Time")
axs[0, 0].legend()
axs[0, 0].grid(True)

#3 SUBPLOT 2: NEUTRON Rate vs Time
axs[0,1].plot(df_muons["time"], df_muons["neutrons"], alpha=0.5, label="Raw Rate")
axs[0,1].plot(df_muons["time"], df_muons["neutrons_ma_24h"],
         linewidth=2,
         label="24h Moving Average")
axs[0,1].xaxis.set_major_locator(mdates.WeekdayLocator(interval=1))
#axs[0,1].set_xlabel("Time",fontsize=8)
axs[0,1].set_ylabel("Neutron Rate".capitalize(),fontsize=12)
axs[0,1].set_title("Neutron Rate vs Time".capitalize())
axs[0,1].legend()
axs[0,1].grid(True)

# Muon to pressure correlation ===============
slope, intercept, r_value, p_value, std_err = linregress(
    df_muons["pressure"],
    df_muons["coincidence"]
)

axs[1,0].scatter(df_muons["pressure"], df_muons["coincidence"], alpha=0.5)
x = "pressure"
y = "coincidence"
sns.regplot(data=df_muons, x=x, y=y, ci=95, scatter=False, line_kws={"color": "red"}, ax=axs[1,0])
axs[1,0].set_title("Muon Rate vs Pressure")
axs[1,0].xaxis.set_major_locator(mdates.WeekdayLocator(interval=1))
axs[1,0].set_xlabel('Pressure (hPa)'.capitalize(), fontsize=12)
axs[1,0].set_ylabel('Muon Rate'.capitalize(), fontsize=12)
axs[1,0].grid(True)
axs[1,0].text(
    0.05,
    0.95,
    f"r = {r_value:.2f}",
    transform=axs[1, 0].transAxes,
    fontsize=12,
    verticalalignment="top",
    bbox=dict(boxstyle="round,pad=0.3", facecolor="white", alpha=0.5),
)

# SUBPLOT 4: Neutron to Pressure Correlation =================

slope, intercept, r_value, p_value, std_err = linregress(
    df_muons["pressure"],
    df_muons["neutrons"]
)

plt.figure(figsize=(8,6))
plt.scatter(df_muons["pressure"], df_muons["neutrons"], alpha=0.5)

x = "pressure"
y = "neutrons"
sns.regplot( data=df_muons, x=x, y=y, ci=95, line_kws={"color": "red", "label": f"Regression Line (r = {r_value:.2f})"})
plt.title('Correlation between Neutron rate and Pressure')
plt.xlabel('Pressure'.capitalize(), fontsize=12)
plt.ylabel('Neutron Rate'.capitalize(), fontsize=12)
plt.grid(True)
plt.gca().text(
    0.05, 0.95, 
    f'r = {r_value:.2f}', 
    transform=plt.gca().transAxes, 
    fontsize=12, 
    verticalalignment='top',
    bbox=dict(boxstyle='round,pad=0.3', facecolor='white', alpha=0.5)
)

# SUBPLOT 4 Muon to Solar Flux =================
# First calculate solar flux based on sunspot data
axs[1,1].plot(df_spots["Date"], df_spots["solarflux"])
axs[1,1].xaxis.set_major_locator(mdates.WeekdayLocator(interval=4))
axs[1,1].xaxis.set_major_formatter(mdates.DateFormatter('%Y-%m-%d'))
axs[1,1].set_xlabel("Date".capitalize(),fontsize=1)
axs[1,1].set_ylabel("Estimated Solar Flux".capitalize(), fontsize=12)
axs[1,1].set_title("Estimated Solar Flux by Day", fontsize=12)
axs[1,1].grid(True)

# Muon to pressure correlation ===============
slope, intercept, r_value, p_value, std_err = linregress(
    df_muons["pressure"],
    df_muons["coincidence"]
)

axs[1,0].scatter(df_muons["pressure"], df_muons["coincidence"], alpha=0.5)
x = "pressure"
y = "coincidence"
sns.regplot(data=df_muons, x=x, y=y, ci=95, scatter=False, line_kws={"color": "red"}, ax=axs[1,0])
axs[1,0].set_title("Muon Rate vs Pressure")
axs[1,0].xaxis.set_major_locator(mdates.WeekdayLocator(interval=1))
axs[1,0].set_xlabel('Pressure (hPa)'.capitalize(), fontsize=12)
axs[1,0].set_ylabel('Muon Rate'.capitalize(), fontsize=12)
axs[1,0].grid(True)
axs[1,0].text(
    0.05,
    0.95,
    f"r = {r_value:.2f}",
    transform=axs[1, 0].transAxes,
    fontsize=12,
    verticalalignment="top",
    bbox=dict(boxstyle="round,pad=0.3", facecolor="white", alpha=0.5),
)

# End definitions ====

plt.tight_layout
#fig.autofmt_xdate()
plt.show()

